01 - 数据工程
这一篇做完,磁盘上会多出一个 train.bin 文件。下一篇写的训练代码,打开它就能直接开训。
需要的前置知识:会写 Python,知道什么是数组和字典。不需要懂深度学习,也不需要读过 Transformer 的论文。涉及的概念我都会在用到的地方先解释一遍。
零、开始之前:模型到底吃什么
这一节不写代码,只建立一个心智模型。如果跳过这一节,后面每一步你都会觉得是在凭空做选择题。
0.1 模型不认识文字
神经网络内部全是矩阵乘法和加法,它处理的只能是数字。所以拿到一句中文或英文,第一件事必须是把它变成数字。
变成数字的方式不是随便 编号。我们要的是:把文本切成一小片一小片,每一片在一张对照表里查到一个编号。这张对照表叫词表(vocabulary),切出来的每一片叫一个 token,查到的编号叫 token id。
原文: unhappiness
切片: un | happi | ness
查表: 1523 | 8891 | 2077
0.2 token 不是字,也不是词
初学最容易卡在这。token 是介于字符和单词之间的东西,中文里没有对应的概念,所以一般不翻译。
为什么不按单词切?因为 单词太多了。英文常用词几十万,加上时态变化、复合词、人名地名,词表根本装不下。而且遇到词表外的词只能标记成「未知」,该位置的信息完全丢失。
为什么不按字符切?因为那样序列太长了。unhappiness 有 11 个字符,就要占 11 个位置。模型每次能看的长度是有限的(我们这个模型是 2048),按字符切能装的内容就少得可怜,而且模型还得自己从字母学起单词的概念,学得非常慢。
所以取中间:常用的词保持完整,生僻的词拆成有意义的片段。unhappiness 拆成 un + happi + ness,三个片段都是高频的,都在词表里。这样词表规模可控(我们用 32000),又不会有装不下的词。
三种切法摆在一起对比:
unhappiness 在三种粒度下的切分结果。子词切分是词表规模和序列长度之间的折中,最下面一行的数字就是查表得到的 token id。怎么决定哪些片段进词表,就是第四节要讲的 BPE 算法干的事。
0.3 训练的时候,一次输入长什么样
模型的训练目标只有一个:看着前面的 token,猜下一个 token 是什么。 这叫 next token prediction,也叫自回归。就这么简单,没有别的。
具体到一次输入。假设有一句话被切成了 5 个 token,id 是 [10, 24, 37, 42, 55]:
输入 x = [10, 24, 37, 42]
标签 y = [24, 37, 42, 55]
y 就是 x 整体右移一位。画出来是这样:
模型看到 10,要预测 24;看到 10, 24,要预测 37;以此类推。四个位置同时算,同时算 loss,这就是为什么 Transformer 训练效率高。
模型每个位置输出的是一个长度 32000 的概率分布(词表里每个 token 各自的概率),拿它和真实的下一个 token 算交叉熵,就是 loss。
这里可以顺手记一个非常有用的数:训练刚开始、模型什么都还没学到的时候,它对 32000 个候选是均匀猜的,此时 loss 应该等于 ln(32000) ≈ 10.37。03 篇第一次跑训练时,如果初始 loss 不在 10.4 附近,说明代码有 bug,不用等它跑完。
0.4 所以数据工程要干的事
现在可以说清楚这一篇的任务了:把一大堆网页文本,变成一长串 token id,存成一个文件。
就这一件事。输入是几千万篇文档,输出是一个几十亿长度的整数数组。
"The cat sat..." ┐
"机器学习是..." ├──► [10, 24, 37, 42, 55, 2, 891, 7734, ...] ──► train.bin
"Once upon a..." ┘
训练的时候,从这个长数组里随便挑一个位置,往后取 2048 个数当 x,右移一位取 2048 个数当 y,就是一条训练样本。
0.5 那为什么要五个步骤
既然任务这么简单,直接 tokenize 一下不就完了?中间那些清洗、去重是干嘛的。
因为原始网页数据很脏,而数据的问题没法在后面的环节补救,只能重训。模型结构写错了,loss 不降,两小时就能发现;数据里混了 5% 的重复网页,你得跑完 17 小时才隐约觉得这模型怎么老在复读,然后 68 卡时的钱已经花掉了。
五个步骤各自解决一个具体问题:
| Step | 做什么 | 不做会怎样 |
|---|---|---|
| 1 下载 | 把语料弄到本地 | — |
| 2 清洗 | 扔掉太短的、结构性垃圾、语种不对的 | 模型学一堆导航栏和分隔线 |
| 3 去重 | 扔掉重复和近似重复的文档 | 重复内容被反复加权,模型爱复读 |
| 4 训 tokenizer | 造出那张「片段 ↔ 编号」的对照表 |